SplitWithOverlap
沿指定轴(axis)将输入张量切分为多个输出张量。与标准 Split 不同,该算子允许通过 start_indices 和 end_indices 自定义每个输出块的起始和结束位置,从而支持输出块之间的重叠(Overlap)。
\[\text{对于第 } j \text{ 个输出张量,其在 axis 轴上的第 } k \text{ 个元素对应:}\]
\[Output[j]_{(\dots, k, \dots)} = Input_{(\dots, start\_indices[j] + k, \dots)} \quad \text{其中 } 0 \le k < (end\_indices[j] - start\_indices[j])\]
- 输入:
input - 输入张量数据地址。
- params - 参数数组。
axis - 进行切分的轴索引。
input_shape - 输入张量的形状数组。
input_ndim - 输入张量的维度。
num_split - 输出张量的数量。
start_indices - 每个输出张量在切分轴上的起始索引数组。
end_indices - 每个输出张量在切分轴上的结束索引数组。
core_mask(int, 可选) - 核掩码(仅适用于共享存储版本)。
- 输出:
outputs - 各个输出张量中填充了切分后的数据。
- 支持平台:
FT78NEMT7004
备注
FT78NE 支持 int8, int16, int32, fp32, fp64, cplx64, cplx128
MT7004 支持 fp16, fp32, int16, int32, cplx64
算子支持不连续切分或有重叠的切分。
每个输出张量在非切分轴上的维度与输入张量保持一致。
共享存储版本:
-
void i8_split_with_overlap_s(int8_t *input, int8_t *outputs[], long long *params, int core_mask)
-
void i16_split_with_overlap_s(int16_t *input, int16_t *outputs[], long long *params, int core_mask)
-
void i32_split_with_overlap_s(int32_t *input, int32_t *outputs[], long long *params, int core_mask)
-
void hp_split_with_overlap_s(half *input, half *outputs[], long long *params, int core_mask)
-
void fp_split_with_overlap_s(float *input, float *outputs[], long long *params, int core_mask)
-
void dp_split_with_overlap_s(double *input, double *outputs[], long long *params, int core_mask)
-
void c64_split_with_overlap_s(float *input, float *outputs[], long long *params, int core_mask)
-
void c128_split_with_overlap_s(double *input, double *outputs[], long long *params, int core_mask)
C调用示例:
1//FT78NE示例(共享存储) 2#include "78NE/utils.h" 3 4int main() { 5 int core_id = get_core_id(); 6 int logic_core_id = GetLogicCoreId(core_mask, core_id); 7 int core_num = GetCoreNum(core_mask); 8 9 // --- 计算尺寸 --- 10 int input_size = input_shape[0] * input_shape[1] * input_shape[2]; 11 int output_sizes[num_split]; 12 int total_output_size = 0; 13 int i; 14 for (i = 0; i < num_split; ++i) { 15 output_sizes[i] = input_shape[0] * (end_indices[i] - start_indices[i]) * input_shape[2]; 16 total_output_size += output_sizes[i]; 17 } 18 19 // --- 内存布局 (DDR空间 - 共享) --- 20 float* input = (float*)0x81000000; 21 float* output_base = (float*)0x82000000; 22 float* check_base = (float*)0x83000000; 23 24 // --- 构建指针数组 (栈空间 - 独享) --- 25 float* outputs[num_split]; 26 float* check_outputs[num_split]; 27 int current_offset = 0; 28 for (i = 0; i < num_split; ++i) { 29 outputs[i] = output_base + current_offset; 30 check_outputs[i] = check_base + current_offset; 31 current_offset += output_sizes[i]; 32 } 33 34 // --- 数据初始化 (DDR空间 - 共享) --- 35 // 根据《多核编程规定》,共享空间变量的并发修改只能发生在单独分支 36 if (logic_core_id == 0) { 37 printf("Test_SplitWithOverlap_multicore_DDR_fp32\n"); 38 for (i = 0; i < input_size; ++i) { 39 input[i] = rand() / (float)RAND_MAX * 100 - 50; 40 } 41 // 清零输出和校验区,避免旧数据干扰 42 memset(output_base, 0, total_output_size * sizeof(float)); 43 memset(check_base, 0, total_output_size * sizeof(float)); 44 } 45 46 long long params[10]; 47 params[0] = (long long) axis; 48 params[1] = (long long) input_shape; 49 params[2] = (long long) input_ndim; 50 params[3] = (long long) num_split; 51 params[4] = (long long) start_indices; 52 params[5] = (long long) end_indices; 53 54 sys_bar(0, core_num); 55 fp_split_with_overlap_s(input, outputs, params, core_mask); 56 sys_bar(0, core_num); 57 return 0; 58}
私有存储版本:
-
void i8_split_with_overlap_p(int8_t *input, int8_t *outputs[], long long *params)
-
void i16_split_with_overlap_p(int16_t *input, int16_t *outputs[], long long *params)
-
void i32_split_with_overlap_p(int32_t *input, int32_t *outputs[], long long *params)
-
void hp_split_with_overlap_p(half *input, half *outputs[], long long *params)
-
void fp_split_with_overlap_p(float *input, float *outputs[], long long *params)
-
void dp_split_with_overlap_p(double *input, double *outputs[], long long *params)
-
void c64_split_with_overlap_p(float *input, float *outputs[], long long *params)
-
void c128_split_with_overlap_p(double *input, double *outputs[], long long *params)
C调用示例:
1//MT7004 示例(私有存储) 2#include <stdio.h> 3 4int main() { 5 int input_size = input_shape[0] * input_shape[1] * input_shape[2]; 6 int output_sizes[num_split]; 7 int total_output_size = 0; 8 int i; 9 for (i = 0; i < num_split; ++i) { 10 output_sizes[i] = input_shape[0] * (end_indices[i] - start_indices[i]) * input_shape[2]; 11 total_output_size += output_sizes[i]; 12 } 13 14 int32_t* input = (int32_t*)0x10010000; 15 int32_t* output_base = (int32_t*)0x10020000; 16 int32_t* check_base = (int32_t*)0x10030000; 17 18 int32_t* outputs[num_split]; 19 int32_t* check_outputs[num_split]; 20 int current_offset = 0; 21 for (i = 0; i < num_split; ++i) { 22 outputs[i] = output_base + current_offset; 23 check_outputs[i] = check_base + current_offset; 24 current_offset += output_sizes[i]; 25 } 26 27 for (i = 0; i < input_size; ++i) { 28 input[i] = (int32_t)i + 0.1f; 29 } 30 memset(output_base, 0, total_output_size * sizeof(int32_t)); 31 memset(check_base, 0, total_output_size * sizeof(int32_t)); 32 33 long long params[10]; 34 params[0] = (long long)axis; 35 params[1] = (long long)input_shape; 36 params[2] = input_ndim; 37 params[3] = num_split; 38 params[4] = (long long)start_indices; 39 params[5] = (long long)end_indices; 40 41 i32_split_with_overlap_p(input, outputs, params); 42 return 0; 43}